Skip to main content

02 - 模型结构从零实现

这一篇产出一个 model.py。跑完验收脚本,它的参数量要正好是 502,193,664,前向能出 logits,初始 loss 落在 10.37 附近。

不用 from_pretrained,不用 transformers 的任何模型类。整个文件只依赖 torchtorch.nn,四百行以内。

需要的前置:知道矩阵乘法,写过一点 PyTorch(会 nn.Linearforward 就够)。01 篇的 train.bin 这一篇用不上,纯写模型。

零、开始之前:Transformer 到底在算什么

0.1 任务还是 01 篇那个任务

再确认一遍目标,因为整个模型结构都是围着它设计的:看着前面的 token,猜下一个 token

输入是一串 token id,输出是每个位置上「下一个 token 是词表里哪个」的概率分布。就这样。

0.2 数据在模型里走一遍

先不管内部细节,看形状怎么变。设 batch 大小为 B、序列长度为 T:

步骤张量形状在干什么
输入(B, T)整数,每个数是一个 token id
Embedding 查表(B, T, 1536)每个 id 换成一个 1536 维向量
Block × 18(B, T, 1536)形状不变,内容被反复加工 18 次
最后的 RMSNorm(B, T, 1536)归一化
lm_head(B, T, 32000)投影到词表大小,得到 logits

画成图:

输入 token id(B, T)整数,01 篇的产出Embedding 查表(B, T, 1536)32000 × 1536,占 9.8% 参数× 18 层,形状始终不变RMSNormGQA Attentiontoken 之间交换信息RMSNormSwiGLU MLP每个 token 各自加工Attention 占 25% 参数MLP 占 75% 参数(B, T, 1536)最后的 RMSNorm(B, T, 1536)lm_head 投影到词表(B, T, 32000)与 Embedding 共享权重
右侧一列是张量形状。中间 18 层的输入输出形状完全相同,都是 (B, T, 1536),这就是 Transformer 能随意堆深的结构性原因。最上和最下的 Embedding 与 lm_head 共享同一组权重(6.3 节)。

关键在于中间 18 层形状完全不变。每一层的输入输出都是 (B, T, 1536),所以可以随便堆几层。这是 Transformer 能做深的结构性原因。

最后那个 (B, T, 32000) 叫 logits,每个位置一个长度 32000 的向量,softmax 之后就是概率分布。

0.3 一个 Block 里有什么

每层 Block 干两件事,顺序固定:

一层 Block:两条残差分支,进出形状完全一样残差分支 ① x = x + Attention(Norm(x))残差分支 ② x = x + MLP(Norm(x))输入 x(B, T, 1536)RMSNormAttentionGQA · 12 Q 头 / 4 KV 头+RMSNormSwiGLU MLP1536 → 4096 → 1536+输出(B, T, 1536)token 之间交换信息第 5 个位置想知道第 2 个位置说了什么,靠它每个 token 各自加工对每个位置独立操作,位置之间不通信残差是恒等通路:从第 18 层到第 1 层存在一条完全没有变换的路,梯度能无损流回去。没有它,18 层的梯度传到第一层基本就没了。Pre-norm 的代价:每层输出方差随层数累加,所以整个模型的最后还要再补一个 RMSNorm,把进 lm_head 之前的数值拉回来。
两个子层的分工是这一层结构的全部内容:Attention 横向打通位置,MLP 纵向加工特征,两者交替。注意 Norm 都在分支里而不在主干上 —— 这个位置差别就是 2.5 节 Pre-norm 与 Post-norm 的区别,它决定了 18 层能不能训得动。

分工可以这么理解。Attention 负责 token 之间的信息交换:第 5 个位置想知道第 2 个位置说了什么,靠它。MLP 负责每个 token 自己的加工:拿到信息之后做非线性变换,它对每个位置独立操作,位置之间不通信。

那两个圆圈是残差连接,也就是 x = x + f(x)。它的作用是给梯度留一条直通的路。没有它,18 层的梯度传到第一层基本就没了。这个结构 02 篇不展开推导,记住「残差是恒等通路」就够用。

0.4 这一篇要写的五个部件

部件作用在哪一节
RMSNorm归一化,稳住数值第二节
RoPE告诉模型 token 的位置第三节
GQA Attentiontoken 之间交换信息第四节
SwiGLU每个 token 自己加工第五节
Block / LLM把上面四个拼起来第六节

每一节的套路都一样:先说这个部件解决什么问题、不要它会怎样,再讲它怎么做,最后给代码。

一、为什么照抄 Llama,改了 GPT-2 的哪五处

不做架构创新,结构照抄 Llama。理由很实在:出了问题可以直接跟现成实现对照排查,而且将来想加载别人的权重也方便。

相对于最经典的 GPT-2,Llama 改了五处。这五处正好就是第二到第六节的内容:

位置GPT-2Llama(我们用的)换掉的理由
归一化LayerNormRMSNorm少算一个均值,快 7% 左右,效果不掉
位置编码可学习的绝对位置RoPE能表达相对位置,且能外推到更长序列
注意力MHAGQAKV cache 直接小 3 倍
FFNGELU,中间维 4dSwiGLU,中间维 8/3 d同参数量下效果更好
norm 位置Post-normPre-norm深层训练稳定得多

另外所有线性层都不带 bias。原因在 2.3 节顺带说。

不做架构创新,照抄 Llama —— 相对 GPT-2 一共动了六处位置GPT-2Llama(我们用的)换掉的理由归一化LayerNormRMSNorm → 第二节少算一个均值,快 7% 左右,效果不掉位置编码可学习的绝对位置表RoPE → 第三节能表达相对位置,且能外推到更长序列注意力MHA · 12 组 KVGQA · 4 组 KV → 第四节KV cache 直接小 3 倍前馈层GELU,中间维 4dSwiGLU,8/3 d → 第五节同参数量下效果更好norm 位置Post-normPre-norm → 2.5 节深层训练稳定得多线性层 bias带 bias全部 bias=False → 2.3 节实测差别很小,少一组参数少一份显存这六行正好就是第二到第六节的内容 —— 换句话说,把 GPT-2 改成 Llama,就是这一篇要动手写的全部。
值得注意的是这六处没有一处是「新想法」,全是把某个部件换成更省或更稳的版本。0.5B 这个规模上,架构创新的收益远小于把已知有效的配置抄对 —— 而且抄对了才有对照物:任何一处行为异常,都能拿现成的 Llama 实现逐层比。

二、RMSNorm

2.1 为什么需要归一化

神经网络堆深了会有个麻烦:每一层的输出分布会漂。第一层输出的数值范围可能是 ±1,传到第十层可能变成 ±100,再往后可能溢出,或者反过来缩到接近 0 梯度消失。

归一化就是在每层入口把数值拉回一个稳定的范围,让后面的层总是面对差不多尺度的输入。

2.2 LayerNorm 在做什么

LayerNorm 对每个 token 的 1536 维向量,做四件事:

LayerNorm(x)=xμσ2+ϵγ+β\text{LayerNorm}(x) = \frac{x - \mu}{\sqrt{\sigma^2 + \epsilon}} \cdot \gamma + \beta

减均值 μ\mu、除标准差 σ\sigma、乘一个可学习的缩放 γ\gamma、加一个可学习的偏置 β\beta

注意这里是对每个 token 自己的 1536 维算均值和方差,不是跨 batch 也不是跨序列。所以它跟 batch size 无关,这点比 BatchNorm 好用。

2.3 RMSNorm 砍掉了什么

RMSNorm 的做法是把「减均值」和「加偏置」都去掉,只留缩放:

RMSNorm(x)=x1di=1dxi2+ϵγ\text{RMSNorm}(x) = \frac{x}{\sqrt{\frac{1}{d}\sum_{i=1}^{d} x_i^2 + \epsilon}} \cdot \gamma

分母那一坨就是均方根(root mean square),名字由此而来。

为什么能砍掉减均值这一步,是这一节的关键。RMSNorm 原论文的观察是:LayerNorm 真正起作用的是「把向量缩放到统一尺度」这件事,而不是「把中心移到 0」。做了消融实验,去掉中心化之后效果基本不掉,但省掉了一遍求均值和一遍减法。

砍掉 bias 也是同样的道理,实测加不加差别很小,而少一组参数就少一份显存和一次加法。这也是为什么整个模型的所有 nn.Linear 都写 bias=False

省下来的计算量不算大,但归一化在每层要做两次、18 层就是 36 次,累积起来推理能快百分之几。免费的收益没理由不要。

对每个 token 自己的 1536 维做,跟 batch size 无关LayerNormGPT-2① 减去均值 μ把中心移到 0② 除以标准差 σ缩放到统一尺度③ 乘可学习 γ让模型自己调回来④ 加可学习偏置 β平移RMSNorm我们用这个砍掉除以均方根 RMS保留的就是这一步乘可学习 γ砍掉为什么砍得掉:原论文的消融结论是 —— LayerNorm 真正起作用的是「缩放到统一尺度」,不是「把中心移到 0」。去掉中心化效果基本不掉。同一个道理砍掉了所有 nn.Linear 的 bias。归一化每层做两次、18 层就是 36 次,累积起来推理快百分之几 —— 免费的收益没理由不要。代价藏在实现里:中间统计量必须转 fp32 再算。BF16 只有 8 位尾数,x.pow(2) 之后直接求和会丢精度,是一类极难定位的数值问题。
四步砍成两步,省的计算量其实很小,真正的启发在于那句消融结论:一个被沿用多年的公式里,有一半的操作可以去掉而不掉效果。这类「先做消融再做加法」的判断方式,比记住 RMSNorm 的公式本身更值得带走。

2.4 代码

class RMSNorm(nn.Module):
def __init__(self, dim: int, eps: float = 1e-5):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim)) # 就是公式里的 gamma

def forward(self, x: torch.Tensor) -> torch.Tensor:
# 统计量始终用 fp32 算,否则 BF16 下 x^2 容易损失精度
dtype = x.dtype
x = x.float()
x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
return x.to(dtype) * self.weight

有个容易忽略的细节:中间统计量要转成 fp32 再算。BF16 只有 8 位尾数,x.pow(2) 之后动态范围会被压缩得很厉害,直接在 BF16 上求和容易丢精度。转 fp32 算完再转回来,代价可以忽略,但能避免一类很难查的数值问题。

weight 初始化成全 1,也就是一开始不做任何缩放,让模型自己学。

2.5 Pre-norm 还是 Post-norm

同样是 RMSNorm,放的位置不同,训练难度差很远。

Post-norm(原始 Transformer 论文的做法)是先过子层再归一化:

x = Norm(x + Attention(x))

Pre-norm(现在几乎所有大模型的做法)是先归一化再过子层:

x = x + Attention(Norm(x))

差别在残差通路。Pre-norm 的写法里,x 是一路直接加过去的,从最后一层到第一层存在一条完全没有归一化操作的恒等通路,梯度可以无损地流回去。Post-norm 里每一层的残差都要再过一次 Norm,梯度传递会被反复缩放,层数一深就容易出问题。

代价是 Pre-norm 每层输出的方差会随层数累加,所以最后要额外补一个 RMSNorm(就是 0.2 节表里倒数第二行那个),把进 lm_head 之前的数值拉回来。

同样是 RMSNorm,放在主干上还是放在分支里,决定 18 层训不训得动Post-norm原始论文子层 f+Norm子层 f+Norm子层 f+Norm主干上每隔一层就夹着一个 Norm,梯度回传要反复穿过它们、被反复缩放,层数一深就出问题Pre-norm我们用这个Norm子层 f+Norm子层 f+Norm子层 f+主干(这条横线)上没有任何操作 —— 从最后一层到第一层是一条恒等通路,梯度无损地流回去代价:Pre-norm 每层输出的方差会随层数累加,所以整个模型最后要额外补一个 RMSNorm,把进 lm_head 之前的数值拉回来。
只看公式 Norm(x + f(x))x + f(Norm(x)) 很难感觉到差别,画成通路就一目了然:前者的主干被 Norm 切成了 18 段,后者的主干是通的。深层网络能不能训,很大程度上就取决于反向传播时有没有这样一条不被打断的路。

三、RoPE 旋转位置编码

3.1 Attention 根本看不见位置

这是一个不那么直观但很重要的事实:attention 本身对输入顺序是无感的

原因在于 attention 的计算方式。每个位置的输出是所有位置的 value 的加权和,权重由 query 和 key 的内积决定。这里面没有任何一项跟「第几个位置」有关。把输入序列打乱顺序,输出也只是跟着打乱,内容完全一样。

所以「我打你」和「你打我」在纯 attention 眼里是同一个东西。必须额外把位置信息喂进去。

3.2 早期做法和它的问题

GPT-2 的做法是搞一个可学习的位置向量表,大小是 (max_seq_len, d_model),第 t 个位置就查第 t 行,加到 embedding 上。

两个问题。一是只能处理训练时见过的长度,表只有 1024 行,第 1025 个位置查不到,模型直接不能用。二是它编码的是绝对位置,但语言里真正重要的往往是相对距离,「形容词修饰它后面第一个名词」这种规律跟这个词在全文第几个位置没关系。

先认清一件事:attention 本身对输入顺序完全无感我 打 你你 打 我纯 attention输出是 value 的加权和,权重由内积决定被当成同一个东西内积里没有任何一项跟位置有关所以位置信息必须从外面额外喂进去GPT-2:可学习的绝对位置表形状 (max_seq_len, d_model) = (1024, 1536),第 t 个位置查第 t 行加到 embedding 上① 只能处理训练时见过的长度 —— 第 1025 个位置查不到,模型直接不能用② 编码的是绝对位置,可语言里重要的往往是相对距离RoPE:把位置变成旋转角度不给 embedding 加任何东西,��而是在 attention 内部把 q 和 k 各自按位置旋转① 内积算完只剩位置差 n − m,相对性是推导出来的,不是学出来的② 没有「表」这个东西,位置多大都算得出来,天然可外推「形容词修饰它后面第一个名词」这类规律,跟这个词在全文第几个位置毫无关系 —— 这就是绝对位置编码最根本的错配。
把「attention 看不见位置」当成事实记住还不够,要记住它的来源:输出是所有 value 的加权和,而权重只由内容决定。正因为位置信息是纯粹外加的,怎么加就成了一个可以自由设计的口子,RoPE 走的是「加在内积里而不是加在输入上」这条路。

3.3 RoPE 的想法:把位置变成旋转角度

RoPE 的做法很巧。它不给 embedding 加东西,而是在 attention 内部,把 query 和 key 按位置旋转一个角度

具体做法:把 128 维的 head 向量两两分成 64 组,每组当成一个二维平面上的向量。位置 mm 处的向量,第 ii 组按角度 mθim\theta_i 旋转:

(xiyi)=(cosmθisinmθisinmθicosmθi)(xiyi)\begin{pmatrix} x_i' \\ y_i' \end{pmatrix} = \begin{pmatrix} \cos m\theta_i & -\sin m\theta_i \\ \sin m\theta_i & \cos m\theta_i \end{pmatrix} \begin{pmatrix} x_i \\ y_i \end{pmatrix}

位置越靠后,转得越多。

3.4 为什么旋转能表达相对位置

这是 RoPE 的精髓,值得推一下。

二维旋转矩阵有个性质:R(α)R(β)=R(βα)R(\alpha)^\top R(\beta) = R(\beta - \alpha)。也就是说一个旋转的转置乘另一个旋转,等于「差角」的旋转。

现在看 attention 里的内积。位置 mm 的 query 和位置 nn 的 key,各自旋转之后再做内积:

(R(mθ)q)(R(nθ)k)=qR(mθ)R(nθ)k=qR((nm)θ)k(R(m\theta)q)^\top (R(n\theta)k) = q^\top R(m\theta)^\top R(n\theta) k = q^\top R((n-m)\theta) k

结果里只剩下 nmn-m绝对位置 mmnn 都消掉了,只留下它们的差。

用图看最直观。两个向量各自转到自己的位置角度上,它们之间的夹角只跟位置差有关,整体一起挪不改变夹角:

q 在位置 1,k 在位置 4qk夹角 3θn − m = 4 − 1 = 3q 在位置 5,k 在位置 8qk夹角 3θn − m = 8 − 5 = 3两边的绝对位置完全不同,但夹角一样,所以内积一样 —— 模型只感知到「隔了 3 个」
内积由夹角决定,而旋转后的夹角恰好是 (n−m)·θ。把 q 和 k 整体往后挪同样的距离,夹角不变,内积也不变。图中的 θ 为便于观看取了较大值,实际第 0 组约 57°、第 63 组约 0.007°。

这正是我们想要的:模型感知到的是「这两个 token 隔多远」,而不是「它们分别在第几个位置」。

原理讲完要验。00-pretrain-0.5b/rope_demo.py 是个纯 Python 的验证,固定两个向量,固定相对距离 3,挪动绝对位置:

固定相对距离 n-m=3,挪动绝对位置:
m n n-m 旋转后内积
0 3 3 -1.7349277069
1 4 3 -1.7349277069
2 5 3 -1.7349277069
3 6 3 -1.7349277069
4 7 3 -1.7349277069
5 8 3 -1.7349277069

改变相对距离,内积才跟着变:
n-m 旋转后内积
0 -0.9800000000
1 -1.7064616876
2 -2.0008180700
3 -1.7349277069
4 -1.2924048008
5 -1.2145554743

小数点后 10 位完全一致。相对距离一变,内积立刻跟着变。

3.5 频率怎么定

64 组用同一个角度是不行的,那样只能表达一种尺度的距离。RoPE 给每组配不同的频率:

θi=base2i/d,base=10000\theta_i = \text{base}^{-2i/d}, \quad \text{base} = 10000

第 0 组频率是 1,转得最快,每挪一个位置就转 1 弧度,用来分辨近距离。最后一组频率是 100001=0.000110000^{-1} = 0.0001,转得极慢,几千个位置才转完一圈,用来分辨远距离。

这样 64 组合起来,就能同时表达从「隔 1 个」到「隔几千个」的各种距离。跟时钟的秒针分针时针是一个道理。

同一个位置序列 m = 0…9,四组频率各自转了多少(指针即旋转角)m=0m=3m=6m=9第 0 组 θ = 1第 16 组 θ = 0.1第 32 组 θ = 0.01第 63 组 θ ≈ 1e-4转得最快分辨近距离转得最慢分辨远距离64 组合起来,就能同时表达从「隔 1 个」到「隔几千个」的各种距离 —— 和钟表上秒针、分针、时针共用一个表盘是同一个道理。θ_i = base^(−2i/d),base = 10000。把 base 调大会让所有频率一起变慢,等效于把位置「压缩」,这是长上下文外推最常用的手段之一。
第三行看起来几乎没动,第四行完全看不出变化 —— 这不是画错了,正是设计意图:低频组在几十个位置的尺度上近似不变,所以它承载的是「隔了几百上千个位置」这种粗粒度信息。如果所有组共用一个频率,模型就只能分辨一种尺度的距离。

base 这个数还有个用处:把它调大能让所有频率变慢,等效于把位置「压缩」,这是长上下文外推最常用的手段之一。我们这篇不展开,但要知道 rope_theta 这个超参是干这个的。

3.6 代码

def build_rope_cache(seq_len, head_dim, theta, device, dtype):
"""预计算 cos / sin 表,形状都是 (seq_len, head_dim)"""
idx = torch.arange(0, head_dim, 2, device=device).float() / head_dim
inv_freq = 1.0 / (theta ** idx) # (head_dim/2,)
pos = torch.arange(seq_len, device=device).float() # (seq_len,)
freqs = torch.outer(pos, inv_freq) # (seq_len, head_dim/2)
emb = torch.cat((freqs, freqs), dim=-1) # (seq_len, head_dim)
return emb.cos().to(dtype), emb.sin().to(dtype)

def rotate_half(x):
half = x.shape[-1] // 2
x1, x2 = x[..., :half], x[..., half:]
return torch.cat((-x2, x1), dim=-1)

def apply_rope(q, k, cos, sin):
cos = cos[None, None, :, :] # 广播到 (B, n_heads, T, head_dim)
sin = sin[None, None, :, :]
return q * cos + rotate_half(q) * sin, k * cos + rotate_half(k) * sin

cossin 只跟位置有关,跟输入内容无关,所以建模型的时候算一次存下来就行,不用每次前向重算。代码里用 register_buffer 挂在模型上,它不是参数、不参与梯度,但会跟着 .to(device) 一起搬到 GPU。

3.7 一个必须知道的坑:两套约定

rotate_half 的写法是把 128 维切成前 64 和后 64,配对方式是「第 0 维配第 64 维」。这是 HuggingFace 的约定。

RoPE 原论文用的是另一套:相邻两维配对,也就是「第 0 维配第 1 维」。

两套都能正常训练,效果也一样,但它们互不兼容。 用 A 约定训出来的权重,拿 B 约定的代码加载,前向输出全是乱的,而且不会报任何错,只会表现为模型胡说八道。

这个坑在 08 篇会真的碰上,因为那时候要把权重喂给自制推理框架。训练端和推理端必须用同一套约定,现在选了 rotate_half,到时候推理框架也得是 rotate_half

同一个 128 维向量,两种配对方式(图中只画前 8 维示意)维 0维 1维 2维 3维 64维 65维 66维 67rotate_halfHuggingFace前 64 维与后 64 维配对第 0 维配第 64 维维 0维 1维 2维 3维 4维 5维 6维 7原论文RoFormer相邻两维配对第 0 维配第 1 维两套都能正常训练,效果也一样,但互不兼容。用 A 训出来的权重拿 B 的代码加载,前向输出全是乱的,而且不报任何错,只表现为模型胡说八道。08 篇把权重喂给自制推理框架时会真碰上这个坑。训练端选了 rotate_half,推理端就必须也是 rotate_half —— 这类约定要在写第一行代码时记下来。
这是整篇里最容易造成「训得好好的,一部署就废」的一处。它的危险性不在于难懂,而在于两边各自都自洽:训练端能收敛,推理端能跑完,唯独拼在一起是错的。08 篇的逐层比对就是专门为这类问题设计的排查手段。

四、GQA 分组查询注意力

4.1 先说 MHA

标准多头注意力(MHA)的做法:把 1536 维切成 12 个头,每个头 128 维,各自独立算一遍 attention,最后拼回来。

每个头都有自己的 Q、K、V 三个投影。多个头的意义是让不同的头关注不同类型的关系,有的头管语法,有的头管指代。

4.2 推理时的真正瓶颈是 KV cache

训练的时候一次算整个序列,MHA 没什么问题。问题出在推理。

推理是一个 token 一个 token 往外吐的。生成第 100 个 token 时,需要用到前面 99 个 token 的 K 和 V。如果每次都重算,复杂度是平方级的,慢得没法用。

所以实际做法是把算过的 K 和 V 缓存起来,这就是 KV cache。每生成一个 token,就往缓存里追加一份 K 和 V。

画出来是这样,每生成一个 token,缓存就长一格:

生成第 1 个 token缓存 1 份 KV新算的生成第 2 个复用前面的,只新算 1 份生成第 3 个生成第 4 个生成第 2048 个2048 份 KV单条序列跑满 2048:GQA 75.5 MB  MHA 226.5 MB服务 100 个并发用户就要乘 100:GQA 7.5 GB  MHA 22.7 GB
缓存长度随生成过程线性增长,而且每个并发用户各要一份。所以在推理服务里 KV cache 经常比模型权重本身还占地方,这也是为什么 KV 头数值得省。

这个缓存很占显存,而且跟并发数成正比。服务 100 个用户就要 100 份。vLLM 的 PagedAttention 解决的就是这块的碎片问题(见 PagedAttention)。

4.3 MQA 和 GQA

既然 KV cache 是瓶颈,一个自然的想法是:能不能让多个头共享同一份 K 和 V

MQA(Multi-Query Attention)走到极端,所有 Q 头共享唯一一组 K/V。缓存直接小 12 倍,但效果掉得比较明显。

GQA(Grouped-Query Attention)取中间:把 Q 头分组,每组共享一组 K/V。我们的配置是 12 个 Q 头、4 个 KV 头,也就是每 3 个 Q 头共享 1 组 K/V。

MHA:每个 Q 头配一个 KV 头Q0Q1Q2Q3…共 12 个KV0KV1KV2KV3…共 12 个KV cache 108.0 KB / tokenGQA:每 3 个 Q 头共享一个 KV 头(我们用这个)Q0Q1Q2Q3Q4Q5…共 12 个,分 4 组KV0KV1…共 4 个KV cache 36.0 KB / token(省 3 倍)MQA:所有 Q 头共享唯一一个 KV 头Q0Q1Q2Q3…共 12 个KV0KV cache 9.0 KB / token,但效果掉得明显Q 头数量三种方案都是 12,所以表达能力差别不大;变的只是 KV 头数,而 KV cache 只跟 KV 头数有关。
三种方案的 Q 头都是 12 个,区别只在下面一排 KV 头的数量。GQA 是折中:KV cache 降到 MHA 的三分之一,而 Q 头一个没少,表达能力基本不受影响。这是现在从 Llama 2 70B 到 Qwen 系列的通行选择。

4.4 我们这个配置省了多少

每个 token 要缓存的字节数:

2×nlayers×nkv_heads×dhead×bytes2 \times n_{\text{layers}} \times n_{\text{kv\_heads}} \times d_{\text{head}} \times \text{bytes}

开头的 2 是因为 K 和 V 各存一份。代进去(BF16 所以每个数 2 字节):

方案n_kv_heads每 token2048 长的一条序列
MHA12108.0 KB226.5 MB
GQA(我们的)436.0 KB75.5 MB

省了整整 3 倍。 这个数字是 00-pretrain-0.5b/count_params.py 算出来的。

同时注意 4.1 节说的:Q 头还是 12 个,所以模型的表达能力基本没损失,掉的只是 K/V 的多样性。这是 GQA 的性价比所在,现在从 Llama 2 70B 到 Qwen 系列基本都用它。

顺带看参数量。因为 K/V 投影变窄了,Attention 的参数也跟着少了:

投影形状参数量
q_proj1536 → 12×128 = 15362,359,296
k_proj1536 → 4×128 = 512786,432
v_proj1536 → 4×128 = 512786,432
o_proj1536 → 15362,359,296
合计6,291,456

如果是 MHA,k/v 各是 2,359,296,Attention 合计会变成 9,437,184。

每 token 的 KV cache = 2 × 18 层 × n_kv_heads × 128 × 2 字节MHA · 12 KV 头一条 2048 长的序列 = 226.5 MB108.0 KB / tokenGQA · 4 KV 头一条 2048 长的序列 = 75.5 MB36.0 KB / tokenK / V 投影变窄,Attention 的参数量跟着少MHA9,437,184GQA6,291,456拆开看:q_proj 2,359,296 · k_proj 786,432 · v_proj 786,432 · o_proj 2,359,296 —— 变窄的只有 k 和 v 两个Q 头仍然是 12 个,模型的表达能力基本没损失,掉的只是 K / V 的多样性 —— 这就是 GQA 的性价比所在。
为什么不干脆用 MQA(所有 Q 头共享一组 K/V,缓存再小 3 倍)?因为效果掉得明显。GQA 是在这条曲线上取的折中点,从 Llama 2 70B 到 Qwen 系列基本都停在这里。注意收益只体现在推理:训练时一次算整个序列,根本不存在 KV cache。

4.5 代码

Q 头和 KV 头数量对不上,算 attention 之前要把 KV 复制对齐:

def repeat_kv(x, n_rep):
"""(B, n_kv_heads, T, hd) -> (B, n_kv_heads * n_rep, T, hd)"""
b, n_kv, t, hd = x.shape
if n_rep == 1:
return x
return (x[:, :, None, :, :]
.expand(b, n_kv, n_rep, t, hd)
.reshape(b, n_kv * n_rep, t, hd))

这里用 expand 而不是 repeat,因为 expand 不真的复制数据,只是改 stride。真正的 KV cache 里只存 4 份,复制只发生在计算的那一刻。 要是这里写成 repeat,显存收益就没了。

完整的 Attention:

class Attention(nn.Module):
def __init__(self, c: Config):
super().__init__()
self.n_heads, self.n_kv_heads = c.n_heads, c.n_kv_heads
self.head_dim = c.head_dim
self.n_rep = c.n_heads // c.n_kv_heads # 12 // 4 = 3

self.q_proj = nn.Linear(c.d_model, c.n_heads * c.head_dim, bias=False)
self.k_proj = nn.Linear(c.d_model, c.n_kv_heads * c.head_dim, bias=False)
self.v_proj = nn.Linear(c.d_model, c.n_kv_heads * c.head_dim, bias=False)
self.o_proj = nn.Linear(c.n_heads * c.head_dim, c.d_model, bias=False)

def forward(self, x, cos, sin):
B, T, _ = x.shape

q = self.q_proj(x).view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
k = self.k_proj(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2)
v = self.v_proj(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2)

q, k = apply_rope(q, k, cos[:T], sin[:T]) # RoPE 只作用于 q 和 k
k, v = repeat_kv(k, self.n_rep), repeat_kv(v, self.n_rep)

out = F.scaled_dot_product_attention(q, k, v, is_causal=True)

out = out.transpose(1, 2).contiguous().view(B, T, -1)
return self.o_proj(out)

两个要点。

RoPE 只加在 q 和 k 上,不加在 v 上。 因为位置信息是通过内积起作用的(3.4 节那个推导),而 v 不参与内积,它只是被加权求和。给 v 加旋转没有意义。

is_causal=True 自带因果掩码。 因果的意思是第 t 个位置只能看到 0 到 t,不能看到未来。画出来就是一个下三角:

注意力权重矩阵(行 = 谁在看,列 = 看谁)t0t1t2t3t4t0××××t1×××t2××t3×t4可见:算注意力时参与屏蔽:分数设为 −∞,softmax 后为 0t0 只能看自己,t4 能看全部前面的。这就是为什么 01 篇 0.3 节里四个位置能同时算 loss 而不作弊。序列 2048 时,显式 mask 矩阵自己就要 8 MB,所以用 is_causal 走 FlashAttention 融合实现。
因果掩码保证第 t 行只有前 t+1 列可见。02 篇 8.4 节的因果性测试,验的就是这张图:改动第 t 列之后的输入,第 t 行的输出必须一点都不变。

F.scaled_dot_product_attention 会自动生成这个下三角掩码,而且走 FlashAttention 的融合实现,比自己建一个 (T, T) 的矩阵又快又省显存。

4.6 顺带说 head_dim 为什么是 128

head_dim = d_model / n_heads = 1536 / 12 = 128

128 是个几乎所有主流模型都在用的数。原因是 FlashAttention 这类融合 kernel 是按 64 / 128 这些尺寸做了专门优化的,取别的值会掉到通用实现上,慢不少。定超参的时候,让 head_dim 落在 64 或 128 是条实用的约束。

五、SwiGLU

5.1 标准 FFN 长什么样

原始 Transformer 的前馈层就两层线性加一个激活:

def ffn(x):
return W2(gelu(W1(x))) # 1536 -> 4096 -> 1536

先升维再降维,中间那层一般取 4 倍宽。它对每个位置独立操作,是模型存储「知识」的主要地方。

5.2 门控的想法

SwiGLU 的改动是加一条门控支路:

SwiGLU(x)=Wdown(SiLU(Wgatex)Wupx)\text{SwiGLU}(x) = W_{\text{down}}\big(\text{SiLU}(W_{\text{gate}}\,x) \odot W_{\text{up}}\,x\big)

\odot 是逐元素相乘。直观理解:up 那一路算出候选值,gate 那一路算出「每一维该放行多少」,两者相乘。SiLU 的输出可以接近 0,等于把某些维度关掉。

相比固定的激活函数,门控让「哪些信息通过」变成数据相关的,模型能学得更细。

标准 FFN:一条直路SwiGLU:多一条门控支路xW1 · 1536→4096GELUW2 · 4096→1536两个矩阵 2 × d × 4d = 8d²激活函数是固定的,对每一维一视同仁xgate_projup_projSiLU×down_proj三个矩阵 3 × d × (8/3)d = 8d² 参数量与左边完全相同up 那一路算出候选值,gate 那一路算出「每一维该放行多少」,两者逐元素相乘。SiLU 的输出可以接近 0,等于把某些维度整个关掉。关键差别只有一句:相比固定的激活函数,门控让「哪些信息�通过」变成了数据相关的,同一个维度对不同输入可以开也可以关。论文《GLU Variants Improve Transformer》只给了消融结果 —— 同参数量下稳定地好一点,没有理论解释,作者原话把它归功于运气。
「门控」这个词容易被当成玄学,画出来其实很朴素:两条支路走同一个输入,一条产内容,一条产开关。真正需要留意的是它带来的第三个矩阵 —— 这正是下一节要把中间维从 4d 调到 8/3 d 的原因。

这类改动 Noam Shazeer 那篇论文(GLU Variants Improve Transformer)做了一组消融,结论是同参数量下 SwiGLU 稳定地好一点。没有特别深刻的理论解释,论文里那句结论大意是把它归功于运气。工程上照用就是了。

5.3 为什么中间维是 8/3 d,不是 4 d

这是个容易被忽略但很实际的问题。

标准 FFN 两个矩阵,参数量是 2×d×dffn2 \times d \times d_{\text{ffn}}。取 dffn=4dd_{\text{ffn}} = 4d 时是 8d28d^2

SwiGLU 有三个矩阵,参数量是 3×d×dffn3 \times d \times d_{\text{ffn}}。要让参数量保持在 8d28d^2 不变:

3×d×dffn=8d2    dffn=83d3 \times d \times d_{\text{ffn}} = 8d^2 \implies d_{\text{ffn}} = \frac{8}{3}d

所以 83×1536=4096\frac{8}{3} \times 1536 = 4096,正好是整数,也正好是 2 的幂,对齐得很舒服。多一个矩阵就把中间维按比例调窄,总参数量不变,这样跟标准 FFN 比才是公平的。

条形长度 = 参数量。多一个矩阵,就把中间维按比例调窄标准 FFN中间维 4d,两个矩阵W1W28d²SwiGLU 若也用 4d三个矩阵,参数直接多 50%gateupdown12d² 比较不公平SwiGLU 实际配置中间维 8/3 d = 4096gateupdown8d² 与第一行齐平8/3 × 1536 = 4096,正好是整数,也正好是 2 的幂 —— 这个巧合让 Llama 这套配置在实现上对齐得很舒服。
第一行和第三行的右端对齐,这就是「公平比较」的含义:换结构时如果不把参数量拉平,效果变好可能只是因为参数变多了。看论文里的架构对比时,第一件该确认的事就是两边参数量是否一致。

5.4 代码

class SwiGLU(nn.Module):
def __init__(self, c: Config):
super().__init__()
self.gate_proj = nn.Linear(c.d_model, c.ffn_dim, bias=False)
self.up_proj = nn.Linear(c.d_model, c.ffn_dim, bias=False)
self.down_proj = nn.Linear(c.ffn_dim, c.d_model, bias=False)

def forward(self, x):
return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))

六、组装

6.1 Block

按 0.3 节那张图,Pre-norm 加两条残差:

class Block(nn.Module):
def __init__(self, c: Config):
super().__init__()
self.attn_norm = RMSNorm(c.d_model, c.norm_eps)
self.attn = Attention(c)
self.mlp_norm = RMSNorm(c.d_model, c.norm_eps)
self.mlp = SwiGLU(c)

def forward(self, x, cos, sin):
x = x + self.attn(self.attn_norm(x), cos, sin)
x = x + self.mlp(self.mlp_norm(x))
return x

注意 x = x + f(norm(x)),不是 x = f(norm(x))。少写那个 x + 就没有残差了,模型能跑但训不深。

6.2 完整模型

class LLM(nn.Module):
def __init__(self, c: Config):
super().__init__()
self.config = c
self.embed = nn.Embedding(c.vocab_size, c.d_model)
self.blocks = nn.ModuleList(Block(c) for _ in range(c.n_layers))
self.norm = RMSNorm(c.d_model, c.norm_eps)
self.lm_head = nn.Linear(c.d_model, c.vocab_size, bias=False)

if c.tie_embeddings:
self.lm_head.weight = self.embed.weight

cos, sin = build_rope_cache(c.max_seq_len, c.head_dim, c.rope_theta,
device="cpu", dtype=torch.float32)
self.register_buffer("cos", cos, persistent=False)
self.register_buffer("sin", sin, persistent=False)

def forward(self, idx, targets=None):
B, T = idx.shape
x = self.embed(idx)
cos, sin = self.cos[:T].to(x.dtype), self.sin[:T].to(x.dtype)
for blk in self.blocks:
x = blk(x, cos, sin)
x = self.norm(x)
logits = self.lm_head(x)

if targets is None:
return logits, None
loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.reshape(-1))
return logits, loss

persistent=False 的意思是这两个 buffer 不存进 checkpoint。它们是纯计算出来的,加载时重新算一遍就行,没必要占 checkpoint 的体积。

6.3 权重共享

self.lm_head.weight = self.embed.weight 这一行让输入的 embedding 表和输出的投影矩阵用同一组权重。

两者形状正好都是 (32000, 1536),物理上可以共享。直觉上也说得通:embedding 是「token id 到语义向量」,lm_head 是「语义向量到 token id」,互为逆向,共享一组参数是合理的。

收益是省 49,152,000 个参数,占总量 9.8%。对小模型来说这个比例不小,所以小模型基本都开权重共享,大模型反而常常不开,因为那时候 embedding 占比已经很低了。

七、初始化

7.1 为什么不能随便初始化

初始权重全设成 0,所有神经元输出一样、梯度一样,永远学不出差异。设得太大,前向传几层就爆掉。设得太小,信号一路衰减到 0。

通行做法是从均值 0、标准差 0.02 的正态分布里采样。0.02 这个数来自 GPT-2,后来大家一直沿用。

7.2 残差累积和缩放

有个专门针对残差结构的补充处理,值得说清楚。

Pre-norm 的每一层都是 x = x + f(x)。假设每个 ff 的输出方差是 σ2\sigma^2,那么 18 层加下来,方差会累积到大约 18σ218\sigma^2。层数越多,进 lm_head 之前的数值越大。

GPT-2 的对策是:把每个残差分支最后那个投影的初始化标准差按 1/2L1/\sqrt{2L} 缩小(L 是层数,2 是因为每层有 attention 和 mlp 两条残差分支)。

for name, p in self.named_parameters():
if name.endswith(("o_proj.weight", "down_proj.weight")):
nn.init.normal_(p, mean=0.0, std=0.02 / (2 * c.n_layers) ** 0.5)

o_projdown_proj 正好就是 attention 和 mlp 各自的出口。缩放之后,无论堆多少层,累积方差都维持在同一量级。

不做这一步模型也能训,但初期 loss 会更抖,学习率要设得更保守。这是一行代码就能拿到的稳定性,没理由不做。

Pre-norm 每层都是 x = x + f(x),18 层加下来方差会累积累积方差第 1 层第 9 层第 18 层不缩放:≈ 18σ²缩放后:维持 ≈ σ²做法:把每个残差分支最后那个投影(o_proj 和 down_proj,正好是 attention 与 mlp 各自的出口)的初始化标准差按 1/√(2L) 缩小。不做也能训,但初期 loss 更抖、学习率得更保守。
两条曲线的差别不体现在最终能不能收敛,而体现在训练最开始那几百步的稳定性 —— 这恰恰是最容易让人误判「模型写错了」的阶段。一行初始化代码换来的确定性,比事后调学习率便宜得多。

八、验收

写完必须验,四项都过了才进 03 篇。验收脚本是 00-pretrain-0.5b/verify_model.py

cd code && python verify_model.py

8.1 参数量对账

00-pretrain-0.5b/count_params.py 是个纯 Python 的独立实现,不 import torch,按公式逐项算。拿它的结果跟 model.num_params() 对,两边独立算出同一个数才说明没写错。

它的实际输出:

配置: d_model=1536 n_layers=18 n_heads=12 n_kv_heads=4 head_dim=128 ffn=4096
ffn / d_model = 2.6667 (8/3 = 2.6667)

单层 Block 内部
Attention (q/k/v/o) 6,291,456 25.0%
SwiGLU MLP 18,874,368 75.0%
2 × RMSNorm 3,072 0.0%
单层合计 25,168,896

整个模型
Embedding(与输出层共享) 49,152,000 9.8%
18 层 Block 453,040,128 90.2%
最后的 RMSNorm 1,536 0.0%
合计 502,193,664 (502.2 M)

502,193,664,跟 00 篇定的配置对上。

值得注意的是 MLP 占了单层参数的 75%。这是 Transformer 的常态,也是为什么做量化、剪枝、MoE 的时候都优先动 MLP,那里油水最多。

8.2 前向形状

输入 (2, 128) 的整数张量,输出必须是 (2, 128, 32000)

这一项主要抓 view / transpose 写错的问题。GQA 那段的形状变换比较绕,n_headsn_kv_heads 很容易填反,填反了 shape 就对不上。

8.3 初始 loss 必须落在 10.37 附近

这是最便宜也最有用的一项检查。

刚初始化的模型对 32000 个候选是均匀猜的,此时交叉熵就是均匀分布的熵:

loss=ln132000=ln32000=10.3735\text{loss} = -\ln\frac{1}{32000} = \ln 32000 = 10.3735

跑出来的数偏离 10.37 太多,说明有问题:

  • 明显偏高(比如 15、20):初始化标准差太大,或者某处数值爆了
  • 明显偏低(比如 7、8):更危险,说明模型能看到不该看到的信息,八成是因果掩码没生效

第二种情况尤其要警惕,因为它的表现是「loss 降得特别快」,看曲线像是训练特别顺利,实际是模型在抄答案。

刚初始化的模型对 32000 个候选是均匀猜的,所以 loss 必然是 ln 32000 = 10.3735510.37351520① 明显偏低(7、8)模型看到了不该看到的信息八成是因果掩码没生效最危险的一种② 落在 10.37 附近正常这一项是整套验收里最便宜也最有用的③ 明显偏高(15、20)初始化标准差设得太大或者某处数值已经爆了比第一种好查,因为它一眼就不对第一种尤其要警惕:它的表现是「loss 降得特别快」,看曲线像是训练格外顺利,实际上模型在抄答案 —— 8.4 节的因果性测试就是专门抓它的。这项检查的成本是跑一次前向,收益是在花掉任何卡时之前,排除掉整类最难查的错误。
10.3735 这个数之所以有用,是因为它完全由词表大小决定,跟模型写得对不对无关 —— 它是一个外部的、不依赖实现的参照点。任何一处让模型「提前知道答案」的 bug,都会让这个数往下掉,而这类 bug 在 loss 曲线上伪装成好消息。

8.4 因果性

这一项专门抓 8.3 里说的第二种错误,做法很直接:

改动第 t 个位置之后的输入,第 t 个位置的输出必须一点都不变

base, _ = model(idx)
modified = idx.clone()
modified[:, t + 1:] = torch.randint(0, vocab_size, (B, T - t - 1))
after, _ = model(modified)

# 位置 0..t 的输出必须完全一致
assert (base[:, :t + 1] - after[:, :t + 1]).abs().max() < 1e-5
# 位置 t+1.. 的输出必须变了(否则说明输入压根没改动,测试本身失效)
assert (base[:, t + 1:] - after[:, t + 1:]).abs().max() > 1e-3

第二个断言容易被忽略但必须有。只写第一个断言的话,如果测试代码本身有 bug(比如没真的改到输入),断言会平凡地通过,测了等于没测。

8.5 顺带看 KV cache 的账

count_params.py 还会打印 KV cache 的估算,这个数 08 篇接推理框架时要用:

KV cache(BF16)
GQA n_kv_heads=4 36,864 B/token = 36.0 KB 序列 2048 时 75.5 MB/条
MHA n_kv_heads=12 110,592 B/token = 108.0 KB 序列 2048 时 226.5 MB/条
GQA 省了 3.0 倍

以及静态显存,对应 00 篇 3.3 节的 8 GB:

静态显存(BF16 混合精度 + AdamW)
BF16 权重 1.00 GB
BF16 梯度 1.00 GB
FP32 master weights 2.01 GB
AdamW 动量 m 2.01 GB
AdamW 动量 v 2.01 GB
小计 8.04 GB

九、常见问题

现象大概率是什么原因
初始 loss 远低于 10.37因果掩码没生效,检查 is_causal=True 有没有传
初始 loss 远高于 10.37初始化 std 太大;或者 RMSNorm 忘了转 fp32
参数量比预期多 49Mtie_embeddings 没生效,检查那行赋值在不在
参数量对不上但差得不多多半是某个 nn.Linear 忘了写 bias=False
shape 报错在 attention 里n_headsn_kv_heads 填反了,或者 repeat_kv 漏调用
显存比算的多很多repeat_kv 写成了 repeat,真复制了数据,改回 expand
加载别人的权重后输出是乱码RoPE 约定不一致,见 3.7 节
训到几百步 loss 变 NaN残差缩放初始化没做,或者学习率太高(03 篇细说)

这一篇不需要 GPU

model.py 在 CPU 上就能建起来跑前向,四项验收全都能在本地做完。结构验对了再去租卡,别把调 shape 的时间花在按小时计费的机器上。